Softmax优化(二) Online Softmax从V0到Grid Stride
测试平台:RTX 5080 (Blackwell, sm_120),M=N=4096,代码在 cuda_code/softmax/online-softmax
前言:笔记一《Softmax优化(一)从Naive到两遍融合》走完了Safe Softmax三遍扫描到两遍融合的路径,结论是两遍融合在RTX 5080上没有提速——省掉的那次读命中了L2而非DRAM。这篇接着AIInfraGuide的教程,从算法推导出发,把Online Softmax的CUDA Kernel从V0单线程一路实现到V4 Grid Stride,重点放在online递推的并行化与寄存器缓存这两个Safe Softmax没有的新问题。
版本核心优化配置时间(ms)相对V0V0单线程一行1×2561.23基准V1block并行+warp shuffle合并规约256×10.1677.4×V2float4向量化256×10.166持平V3寄存器缓存one-pass256×10.165持平V4grid stride256×10.173持平
V1之后各版本都在~0.17ms附近打平,与笔记一的结论一致:瓶颈是DRAM带宽,V1已经把访存压到位,后续优化省的都是L2命中或寄存器,不省DRAM。
1 Online Softmax的递推与并行化
Safe Softmax要三遍扫描(max→sum→归一化),Online Softmax把max和sum合并到一遍,读次数从3降到2。递推维护running (m,d):
遇到新元素x,设 m_new = max(m, x):
d = d * exp(m - m_new) + exp(x - m_new)
m = m_new
当x<=m时exp(m-m_new)=1,退化为普通累加;当x>m时旧sum乘以exp(m-x)缩放到新基准。
递推本身是串行的,但可以并行化,关键是两个局部结果(m1,d1)、(m2,d2)的合并操作满足结合律:
m_merged = max(m1, m2)
d_merged = d1 * exp(m1 - m_merged) + d2 * exp(m2 - m_merged)
所以可以先让每个线程算自己那段的局部(m,d),再用树形规约合并。这正是V1的基础。
2 V0:单线程Online Softmax
一个线程串行处理一行,直接按上面的递推扫一遍,再归一化:
__global__ void online_softmax_v0_naive(const float* input, float* output, int M, int N) {
int row = blockIdx.y * blockDim.y + threadIdx.y;
if (row >= M) return;
const float* x = input + row * N;
float* y = output + row * N;
// 第1遍:online同时更新(m,d)
float m = -INFINITY, d = 0.0f;
for (int i = 0; i < N; i++) {
float m_new = fmaxf(m, x[i]);
d = d * expf(m - m_new) + expf(x[i] - m_new);
m = m_new;
}
// 第2遍:归一化
float inv_d = 1.0f / d;
for (int i = 0; i < N; i++) y[i] = expf(x[i] - m) * inv_d;
}
结果~1.23ms,约163 GB/s。和Safe Softmax的V0一样,单线程独享一行,并行度为零、非合并访存,只作正确性基线。
3 V1:Block并行+warp shuffle合并规约
一个block处理一行,每个线程先串行算自己那段的局部(m,d),再做两级归约:warp内用__shfl_down_sync合并,warp间经shared memory交换,warp0收尾。
__device__ void warpReduceOnline(float& m, float& d) {
for (int offset = 16; offset > 0; offset >>= 1) {
float m2 = __shfl_down_sync(0xffffffff, m, offset);
float d2 = __shfl_down_sync(0xffffffff, d, offset);
float m_new = fmaxf(m, m2);
d = d * expf(m - m_new) + d2 * expf(m2 - m_new);
m = m_new;
}
}
__global__ void online_softmax_v1(...) {
// 每线程局部(m,d)
float local_m = -INFINITY, local_d = 0.0f;
for (int i = tid; i < N; i += blockDim.x) { ... } // online递推
// warp内合并 + warp间smem交换 + warp0收尾
warpReduceOnline(local_m, local_d);
...
// 第2遍:归一化
}
结果~0.167ms,约1208 GB/s,比V0快7.4倍。
4 V2:float4向量化
算法不变,只优化访存:把float重解释成float4,一次LDG.128读4个、一次STG.128写4个。但online递推仍是逐元素更新(m,d),展开4个分量依次做。
const float4* x4 = reinterpret_cast<const float4*>(x);
float4* y4 = reinterpret_cast<float4*>(y);
int N4 = N / 4;
// 第1遍:一次读4个,逐个更新(m,d)
float4 v = x4[i];
float mn;
mn = fmaxf(local_m, v.x); local_d = local_d * expf(local_m - mn) + expf(v.x - mn); local_m = mn;
// v.y / v.z / v.w同理
结果~0.166ms,与V1持平。访存指令数÷4、循环迭代÷4,但DRAM已经接近满负荷,指令减少只能缓解不能突破带宽上限。
5 V3:寄存器缓存one-pass(编译期展开的坑)
V1/V2都读2遍input,V3想把第1遍读到的元素缓存到寄存器,第2遍直接读寄存器,把全局访存压到1读+1写。
第一版用运行时下标reg_cache[count],结果ptxas显示缓存被放到了stack(local memory)而不是寄存器:
// 运行时下标:128 bytes stack frame, 25 registers
float reg_cache[MAX_ELEMS];
for (int i = tid; i < N; i += blockDim.x) {
reg_cache[count++] = x[i]; // count是运行时变量
}
要真正留在寄存器,下标必须是编译期常量。把每线程元素数ELEMS做成模板参数,循环用#pragma unroll完全展开:
template <int ELEMS>
__global__ void online_softmax_v3_regcache(...) {
float reg_cache[ELEMS];
#pragma unroll
for (int k = 0; k < ELEMS; k++) {
int i = tid + k * blockDim.x;
float xi = (i < N) ? x[i] : -INFINITY;
reg_cache[k] = xi; // k是编译期常量→寄存器
...
}
}
ptxas对比:展开后ELEMS=16时0 bytes stack frame、56 registers,缓存真正进了寄存器(ELEMS=32是88,ELEMS=64是244)。
结果~0.165ms,仍然与V2持平。原因和笔记一一致:V2省不掉的那次读本来就命中L2,V3从L2再压到寄存器收益更小,反而寄存器压力从34涨到56,占用率下降抵消了省读收益。
6 V4:Grid Stride
前面都是一block一行(grid=M),M很大时会启动海量block、调度开销线性增加。V4用固定block数循环覆盖所有行,算法与V2相同,区别只在block与行的映射从1:1变成1:多。
for (int row = blockIdx.x; row < M; row += gridDim.x) {
// 一整行:online递推 + warp shuffle合并 + 归一化,与V2相同
}
共享内存数组移到循环外跨行复用,grid=min(M, sm_count×4)。结果~0.173ms,与V2持平——M=4096时DRAM已打满,调度不是瓶颈,grid stride的收益只在M极大时体现。
7 为什么V1之后全打平
整条路径下来,V1之后全部卡在~0.17ms。结合笔记一的ncu结论:这张卡的softmax是memory-bound,V1已经做到「读2次+写1次」且额外读命中L2,真实DRAM流量早就到了下限「读1次+写1次」,后续的float4、寄存器缓存、grid stride省的都是L2带宽或调度开销,不省DRAM,所以不加速。
8 总结与FlashAttention的联系
单看「写个softmax kernel」,优化招式(合并访存、向量化、规约、寄存器分块、grid stride)都是GEMM那套工具库的复用。Online Softmax真正的新东西是算法层面:running (m,d)递推 + 合并结合律,这才让并行规约成为可能。它的价值不在「省读」,而在FlashAttention里不物化N×N分数矩阵——只能靠online softmax边算边归一化,跨KV block重缩放。
参考:
